07. 回放缓冲区

回放缓冲区

大多数现代强化学习算法都使用一个回放存储器或缓冲区来存储和回调经验元组。

首先,你可以定义一个很小的数据结构来存储每个经验元组:

from collections import namedtuple
Experience = namedtuple("Experience",
    field_names=["state", "action", "reward", "done", "next_state"])

然后你可以定义一个回放缓冲区类,以存储 Experience 对象列表。为了避免耗尽计算机的内存,建议为你的缓冲区定义大小上限,并使用循环缓冲区(旧内存被新内存覆盖)等列表。

尝试在下面的编程练习中实现该功能,即 add() sample() 方法。点击 测试答案 以执行你的代码并点击 提交 以对其运行测试。尝试自己实现该功能,但是如果遇到问题,可以看一眼解决方案 。

Start Quiz:

"""Replay Buffer."""

import random

from collections import namedtuple

Experience = namedtuple("Experience",
    field_names=["state", "action", "reward", "next_state", "done"])


class ReplayBuffer:
    """Fixed-size circular buffer to store experience tuples."""

    def __init__(self, size=1000):
        """Initialize a ReplayBuffer object."""
        self.size = size  # maximum size of buffer
        self.memory = []  # internal memory (list)
        self.idx = 0  # current index into circular buffer
    
    def add(self, state, action, reward, next_state, done):
        """Add a new experience to memory."""
        # TODO: Create an Experience object, add it to memory
        # Note: If memory is full, start overwriting from the beginning
    
    def sample(self, batch_size=64):
        """Randomly sample a batch of experiences from memory."""
        # TODO: Return a list or tuple of Experience objects sampled from memory

    def __len__(self):
        """Return the current size of internal memory."""
        return len(self.memory)


def test_run():
    """Test run ReplayBuffer implementation."""
    buf = ReplayBuffer(10)  # small buffer to test with
    
    # Add some sample data with a known pattern:
    #     state: i, action: 0/1, reward: -1/0/1, next_state: i+1, done: 0/1
    for i in range(15):  # more than maximum size to force overwriting
        buf.add(i, i % 2, i % 3 - 1, i + 1, i % 4)
    
    # Print buffer size and contents
    print("Replay buffer: size =", len(buf))  # maximum size if full
    for i, e in enumerate(buf.memory):
        print(i, e)  # should show circular overwriting
    
    # Randomly sample a batch
    batch = buf.sample(5)
    print("\nRandom batch: size =", len(batch))  # maximum size if full
    for e in batch:
        print(e)
"""Replay Buffer (solution)."""

import random

from collections import namedtuple

Experience = namedtuple("Experience",
    field_names=["state", "action", "reward", "next_state", "done"])


class ReplayBuffer:
    """Fixed-size circular buffer to store experience tuples."""

    def __init__(self, size=1000):
        """Initialize a ReplayBuffer object."""
        self.size = size  # maximum size of buffer
        self.memory = []  # internal memory (list)
        self.idx = 0  # current index into circular buffer
    
    def add(self, state, action, reward, next_state, done):
        """Add a new experience to memory."""
        e = Experience(state, action, reward, next_state, done)
        if len(self.memory) < self.size:
            self.memory.append(e)
        else:
            self.memory[self.idx] = e
            self.idx = (self.idx + 1) % self.size
    
    def sample(self, batch_size=64):
        """Randomly sample a batch of experiences from memory."""
        return random.sample(self.memory, k=batch_size)

    def __len__(self):
        """Return the current size of internal memory."""
        return len(self.memory)


def test_run():
    """Test run ReplayBuffer implementation."""
    buf = ReplayBuffer(10)  # small buffer to test with
    
    # Add some sample data with a known pattern:
    #     state: i, action: 0/1, reward: -1/0/1, next_state: i+1, done: 0/1
    for i in range(15):  # more than maximum size to force overwriting
        buf.add(i, i % 2, i % 3 - 1, i + 1, i % 4)
    
    # Print buffer size and contents
    print("Replay buffer: size =", len(buf))  # maximum size if full
    for i, e in enumerate(buf.memory):
        print(i, e)  # should show circular overwriting
    
    # Randomly sample a batch
    batch = buf.sample(5)
    print("\nRandom batch: size =", len(batch))  # maximum size if full
    for e in batch:
        print(e)

注意 :这个循环缓冲区只是管理有限内存的一种方式。另一种方法是随机覆盖元素。你甚至可以使用优先级值控制每个经验元组被保留/覆盖的概率。在抽样以实现优先经验回放时也可以采用该方法。